如何求过点(a,b)与(x,0)的直线的最大y截距及对应点索引
嘿,我来帮你优化这个问题的解法!你的O(m*n)思路虽然直观,但当m和n都比较大的时候(比如上万级别),肯定会超时。我们可以通过凸包技巧+二分查找把时间复杂度降到O(n log n + m log n),完美解决大规模数据的问题。
问题回顾
我们有两类点:
- 一类是
(a_i, b_i)(记为集合S),所有a_i都是正整数,且存在最大值max_a; - 另一类是
(x_j, 0)(记为集合T),所有x_j都是正整数,且x_j > max_a。
对于每个x_j,我们需要找到所有(a_i, b_i),使得它们与(x_j, 0)构成的直线的y截距最大,返回这些点的索引。
核心分析:截距的数学变形
直线过(x_j, 0)和(a_i, b_i),其y截距的计算公式为:
k_i = (b_i * x_j) / (x_j - a_i)
因为x_j > a_i,分母为正,所以最大化k_i等价于最大化b_i/(x_j - a_i)。
对于两个点P=(a1,b1)和Q=(a2,b2),我们可以计算出一个临界值x0:当x_j > x0时,P的截距更大;当x_j < x0时,Q的截距更大;当x_j = x0时,两者截距相等。x0的计算公式为:
x0 = (b1*a2 - b2*a1) / (b1 - b2) (当b1≠b2时)
如果b1 = b2,那么a更大的点截距始终更大(因为x_j -a更小,b/(x_j -a)更大),所以这类点我们只需要保留a最大的那个。
优化步骤
1. 预处理:过滤冗余点
- 首先,把集合S中的点按
a降序排序,若a相同则保留b最大的点(相同a下,b大的截距更大)。 - 接着,去除被完全支配的点:如果存在点Q,使得
a_Q ≥ a_P且b_Q ≥ b_P,那么P的截距永远不会超过Q,可以直接剔除。
2. 维护候选点的凸壳(单调栈)
我们需要维护一个单调栈,栈中的点满足相邻点的临界值x0严格递增。这样我们就能通过二分查找快速定位每个x_j对应的最优点。
维护栈的规则:
- 初始化空栈。
- 遍历预处理后的点(按
a升序排序):- 当栈中至少有两个点时,取栈顶的两个点
R和Q(Q是栈顶前一个元素),计算x0(Q,R)和x0(R, 当前点P)。 - 如果
x0(Q,R) ≥ x0(R,P),说明R在任何x_j下都不会是最优的,弹出R;重复此过程直到栈中不足两个点或x0(Q,R) < x0(R,P)。 - 将当前点
P压入栈中。
- 当栈中至少有两个点时,取栈顶的两个点
3. 查询每个x_j
对于每个x_j:
- 在栈的相邻点临界值列表中,用二分查找找到第一个大于
x_j的x0,对应的前一个栈元素就是最优点。 - 如果
x_j恰好等于某个x0,则对应的两个栈元素的截距相等,都是最大值,需要同时返回它们的索引。
代码示例(Python)
def preprocess_points(points): # points是列表,每个元素是(a, b, index) # 第一步:按a降序,b降序排序,去重相同a的点,保留b最大的 points.sort(key=lambda x: (-x[0], -x[1])) unique_points = [] prev_a = -1 for p in points: a, b, idx = p if a != prev_a: unique_points.append(p) prev_a = a # 第二步:维护单调栈,构建凸壳 stack = [] x0_list = [] # 存储栈中相邻点的x0值 def calculate_x0(p1, p2): a1, b1, _ = p1 a2, b2, _ = p2 if b1 == b2: # b相等时,a大的点始终更优,这里不会出现这种情况(因为已经去重) return float('inf') numerator = b1 * a2 - b2 * a1 denominator = b1 - b2 return numerator / denominator for p in unique_points: while len(stack) >= 2: q = stack[-2] r = stack[-1] x0_qr = calculate_x0(q, r) x0_rp = calculate_x0(r, p) if x0_qr >= x0_rp: # 弹出r stack.pop() x0_list.pop() else: break if stack: x0 = calculate_x0(stack[-1], p) x0_list.append(x0) stack.append(p) return stack, x0_list def query_optimal_points(x, stack, x0_list): # 二分查找第一个大于x的x0 left = 0 right = len(x0_list) while left < right: mid = (left + right) // 2 if x0_list[mid] > x: right = mid else: left = mid + 1 # 最优点是stack[left] optimal_indices = [stack[left][2]] # 检查是否x等于某个x0,此时前一个点也最优 if left > 0 and abs(x0_list[left-1] - x) < 1e-9: optimal_indices.append(stack[left-1][2]) return optimal_indices # 示例使用 if __name__ == "__main__": # 输入点:(a, b, index) points = [(1,5,0), (2,3,1), (3,2,2)] stack, x0_list = preprocess_points(points) print("栈中的点:", stack) print("x0列表:", x0_list) # 查询x=4 print("x=4的最优索引:", query_optimal_points(4, stack, x0_list)) # 应该返回[2] # 查询x=5 print("x=5的最优索引:", query_optimal_points(5, stack, x0_list)) # 应该返回[0] # 查询x=3.5(虽然题目中x是整数,这里测试临界值) print("x=3.5的最优索引:", query_optimal_points(3.5, stack, x0_list)) # 返回[0,1]
时间复杂度分析
- 预处理阶段:排序是O(n log n),维护栈是O(n)(每个点最多入栈和出栈一次),总预处理时间O(n log n)。
- 查询阶段:每个查询是O(log n)(二分查找),m个查询总时间O(m log n)。
整体时间复杂度O(n log n + m log n),相比O(m*n)有质的提升,适合处理大规模数据。
内容的提问来源于stack exchange,提问作者Snip3r
相关产品推荐
相关产品推荐

